Conv2DBackpropFilterFusion
计算二维卷积对权重的反向梯度(filter / weight gradient)。
x、dy 为 NHWC;支持普通卷积、Depthwise,以及满足条件时的 1x1 优化路径。
对组 \(g\)、输出通道 \(c_o\)、输入通道 \(c_i\)、核偏移 \((k_h, k_w)\),令 \(C_o' = C_{out}/G\),\(C_i' = C_{in}/G\)。对应输入坐标为:
当 \((y, x)\) 落在输入有效范围内时累加,否则该项为 0:
其中 \(s_h/s_w\) 为步长,\(d_h/d_w\) 为膨胀系数。
实现上将 dw 按组展平存储为
\([C_o',\, K_h \cdot K_w \cdot C_i']\),元素总数为
\(C_{out} \cdot K_h \cdot K_w \cdot C_i'\)。
可选 bias 在部分路径上会叠加到矩阵乘结果;无偏置时传 NULL。
- 输入:
input_dy - 输出梯度地址,形状
[batch, out_h, out_w, out_channel]input_x - 前向输入地址,形状
[batch, in_h, in_w, in_channel]bias - 可选偏置;无偏置时传
NULLconv_param -
ConvParameter结构体地址,字段见下core_mask - 核掩码(仅共享存储版本使用)
ConvParameter 定义:
1typedef struct ConvParameter {
2 void* workspace_; // 工作区指针;实际分配字节数见下方警告(常大于 buffer_size_)
3 int output_batch_; // dy 的 batch
4 int input_batch_; // x 的 batch(通常与 output_batch_ 相同)
5 int input_h_; // 输入高
6 int input_w_; // 输入宽
7 int output_h_; // 输出梯度高
8 int output_w_; // 输出梯度宽
9 int input_channel_; // 输入通道
10 int output_channel_; // 输出通道
11 int kernel_h_; // 卷积核高
12 int kernel_w_; // 卷积核宽
13 int group_; // 组数,1 为普通卷积
14 int pad_l_; // 左填充
15 int pad_u_; // 上填充
16 int dilation_h_; // 高方向膨胀
17 int dilation_w_; // 宽方向膨胀
18 int stride_h_; // 高方向步长
19 int stride_w_; // 宽方向步长
20 int buffer_size_; // 分块预算,单位字节(非元素个数);_s/_p 都读此字段。见下方警告。
21 int pad_d_; // 下填充(conv2d 结构体无此字段)
22 int pad_r_; // 右填充(conv2d 结构体无此字段)
23} ConvParameter;
- 输出:
out_dw - 权重梯度地址;调用前通常先清零,再由算子累加写入; 元素数为
out_channel * kernel_h * kernel_w * (in_channel / group)
- 支持平台:
FT78NEMT7004
警告
MT7004的workspace_ / buffer_size_ 分配算法(与前向 conv2d / transpose 都不同):
buffer_size_ 是字节(不是元素个数)。内部用 buffer_size_/2 估半区;部分汇编按 float 个数做指针算术,因此实际分配 workspace_ 须 ≥ buffer_size_ × sizeof(float)。字段值与分配字节 可以不相等。
共享 _s 与私有 _p 都会读 buffer_size_ / workspace_ (不像前向 conv _s 写死 AM、也不像 transpose _s 写死每核 2048)。
共享多核还会在 workspace_ 里放各核 float partial dw,再归约到 out_dw。若只按 buffer 预算分配、忘了 partial,多核会越界。 分配取:
dw_elems = Cout * Kh * Kw * (Cin / G); workspace_bytes = max(buffer_size_ * sizeof(float), core_num * dw_elems * sizeof(float));
buffer_size_ 下限:2 × (Kh × Kw × Cin / G) × sizeof(dtype) (保证 chunk ≥ 1)。
字段填 buffer_size_(如 2048 或测试常用 0x10000),实际分配见上。
备注
FT78NE 支持 fp32
MT7004 支持 fp16、fp32
Depthwise 条件:output_channel、group、input_channel 三者相等,且 dilation 为 1
1x1 优化条件:Kh=Kw=1,四向 pad 为 0,dilation/stride 为 1,group 为 1
共享存储版本:
-
void hp_conv_grad_filter_s(float16 *input_dy, float16 *input_x, float16 *out_dw, float16 *bias, ConvParameter *conv_param, int core_mask)
-
void fp_conv_grad_filter_s(float *input_dy, float *input_x, float *out_dw, float *bias, ConvParameter *conv_param, int core_mask)
C调用示例:
1// MT7004 示例(共享存储多核,DDR 地址)
2void TestConvGradFilterSMCFp32(int core_mask) {
3 int core_id = get_core_id();
4 int logic_core_id = GetLogicCoreId(core_mask, core_id);
5 int core_num = GetCoreNum(core_mask);
6 float *input_dy = (float *)0x86000000;
7 float *input_x = (float *)0x87000000;
8 float *out_dw = (float *)0x88000000;
9 ConvParameter *param = (ConvParameter *)0x89000000;
10 float *workspace = (float *)0x8A000000;
11 if (logic_core_id == 0) {
12 param->workspace_ = workspace;
13 // 字段填字节预算;实际 malloc 见上警告(可大于 buffer_size_)
14 param->buffer_size_ = 0x10000;
15 param->input_batch_ = 4;
16 param->output_batch_ = 4;
17 param->input_h_ = 4;
18 param->input_w_ = 4;
19 param->input_channel_ = 4;
20 param->output_h_ = 2;
21 param->output_w_ = 2;
22 param->output_channel_ = 4;
23 param->kernel_h_ = 3;
24 param->kernel_w_ = 3;
25 param->group_ = 1;
26 param->pad_l_ = 0;
27 param->pad_u_ = 0;
28 param->pad_d_ = 0;
29 param->pad_r_ = 0;
30 param->dilation_h_ = 1;
31 param->dilation_w_ = 1;
32 param->stride_h_ = 1;
33 param->stride_w_ = 1;
34 }
35 sys_bar(0, core_num);
36 fp_conv_grad_filter_s(input_dy, input_x, out_dw, NULL, param, core_mask);
37}
38
39void main() {
40 int core_mask = 0b1111;
41 TestConvGradFilterSMCFp32(core_mask);
42}
私有存储版本:
-
void hp_conv_grad_filter_p(float16 *input_dy, float16 *input_x, float16 *out_dw, float16 *bias, ConvParameter *conv_param)
-
void fp_conv_grad_filter_p(float *input_dy, float *input_x, float *out_dw, float *bias, ConvParameter *conv_param)
C调用示例:
1// MT7004 示例(私有存储单核,AM 地址)
2void TestConvGradFilterAMFp32(void) {
3 float *input_dy = (float *)0x10000000;
4 float *input_x = (float *)0x10010000;
5 float *out_dw = (float *)0x10020000;
6 ConvParameter *param = (ConvParameter *)0x10040000;
7 float *workspace = (float *)0x10050000;
8 param->workspace_ = workspace;
9 // 字段填字节预算;workspace 建议至少 buffer_size_ * sizeof(float)
10 param->buffer_size_ = 0x10000;
11 param->input_batch_ = 4;
12 param->output_batch_ = 4;
13 param->input_h_ = 4;
14 param->input_w_ = 4;
15 param->input_channel_ = 4;
16 param->output_h_ = 2;
17 param->output_w_ = 2;
18 param->output_channel_ = 4;
19 param->kernel_h_ = 3;
20 param->kernel_w_ = 3;
21 param->group_ = 1;
22 param->pad_l_ = 0;
23 param->pad_u_ = 0;
24 param->pad_d_ = 0;
25 param->pad_r_ = 0;
26 param->dilation_h_ = 1;
27 param->dilation_w_ = 1;
28 param->stride_h_ = 1;
29 param->stride_w_ = 1;
30 fp_conv_grad_filter_p(input_dy, input_x, out_dw, NULL, param);
31}
32
33void main() {
34 TestConvGradFilterAMFp32();
35}